04 - 从单卡到多卡
这一篇把 03 篇的 train.py 改成能在多张卡上跑,产出 train_ddp.py。同时要回答一个问题:多卡到底快了多少,代价是什么。
这也是整个专题里唯一必须真开多卡才能验的一篇。原理和账我会算清楚,但最终那张实测表得在租来的机器上填。
前置:03 篇的单卡训练已经跑通。
零、开始之前:多卡在并行什么
0.1 三种切法
一个模型要放到多张卡上,无非是切三样东西:
| 叫法 | 切什么 | 直觉 |
|---|---|---|
| 数据并行 | 切数据 | 每张卡放一份完整模型,各自算不同的数据,再把梯度对齐 |
| 张量并行 | 切单个矩阵 | 一个大矩阵乘法拆到多张卡上算,算完拼起来 |
| 流水线并行 | 切层 | 前 9 层在卡 0,后 9 层在卡 1,像流水线一样传递 |
真正的大模型训练是三种混着用,叫 3D 并行。
0.2 我们只用数据并行
0.5B 的模型单卡放得下(02 篇算过静态显存 8 GB),所以张量并行和流水线并行完全用不上。它们是为了「一张卡装不下一个模型」而存在的,我们没这个问题。
这一篇只讲数据并行,但会讲透:DDP 怎么工作、通信量多大、什么时候它不够用、不够用了换成 FSDP 又要付什么代价。
0.3 数据并行的基本流程
每张卡各拿一批不同的数据,各自跑前向反向,得到各自的梯度。这些梯度不一样,因为看的数据不一样。
关键在于:参数更新前必须让所有卡的梯度变成同一个值,否则各卡的参数就会越 走越远,最后变成四个不同的模型。
对齐的方式是求平均。这个「把所有卡的数据加起来再分回去」的操作,就是下一节的 all-reduce。
一、集合通信
1.1 几个原语
多卡之间怎么交换数据,是一组标准操作,叫集合通信(collective communication)。NVIDIA 的实现叫 NCCL,PyTorch 底层用的就是它。
| 操作 | 干什么 |
|---|---|
| broadcast | 把某张卡的数据复制给所有卡 |
| reduce | 把所有卡的数据求和(或求最大等),结果放在一张卡上 |
| all-reduce | 求和,结果每张卡都有一份 |
| reduce-scatter | 求和,但每张卡只拿走结果的一部分 |
| all-gather | 每张卡的一小块,拼成完整的一份,每张卡都有 |
要记的关键关系是:all-reduce = reduce-scatter + all-gather。这个拆解决定了它的通信量,也是 FSDP 的实现基础。
1.2 ring all-reduce 的通信量
朴素的做法是所有卡都把数据发给卡 0,卡 0 加完再发回去。这样卡 0 的网卡要承担 N 倍流量,成为瓶颈。
实际用的是 ring all-reduce:把 N 张卡连成一个环,数据切成 N 块。
数据切成 N 块,分两阶段各转 N-1 步:
第一阶段 reduce-scatter,转 N-1 步,每步每张卡收发一块,结束后每张卡持有完整结果的其中一块。第二阶段 all-gather,再转 N-1 步,把每块广播到所有卡。
所以每张卡收发的数据量是:
N 很大时这个系数趋近 2,跟卡数无关。这是 ring all-reduce 的核心优点,也是它能扩展到几千张卡的原因。
二、DDP
2.1 原理
PyTorch 的 DistributedDataParallel(DDP)做的事:
- 启动时把 rank 0 的参数 broadcast 给所有卡,保证起点一致
- 每张卡各自前向反向
- 反向过程中,梯度一算出来就 all-reduce,跟反向计算重叠
- 每张卡拿到平均梯度,各自
optimizer.step()
第 3 点是 DDP 快的关键,也是它和朴素实现的主要差距:
注意每张卡都是完整地跑一遍 optimizer.step(),输入(平均梯度)相同、初值相同,所以结果必然相同。参数不需要额外同步。
2.2 代码改动
03 篇的 train.py 改成多卡,实际只动四个地方。
一是初始化进程组、绑定设备。
def setup_dist():
dist.init_process_group(backend="nccl")
rank = int(os.environ["RANK"])
local_rank = int(os.environ["LOCAL_RANK"])
world_size = int(os.environ["WORLD_SIZE"])
torch.cuda.set_device(local_rank) # 必须在建模型之前
return rank, local_rank, world_size, torch.device(f"cuda:{local_rank}")
RANK、LOCAL_RANK、WORLD_SIZE 这三个环境变量由 torchrun 注入,不用自己管。rank 是全局编号,local_rank 是这台机器内的编号,单机情况下两者相同。
torch.cuda.set_device(local_rank) 必须在建模型之前调用,否则所有进程会往卡 0 上堆,直接 OOM。
二是包一层 DDP。
model = DDP(model, device_ids=[local_rank])
三是随机种子要错开。
torch.manual_seed(args.seed + rank)
np.random.seed(args.seed + rank)
我们的 get_batch 是随机取起点的,如果所有卡用同一个种子,四张卡会取到完全一样的数据,等于白算三份。加上 rank 就错开了。
四是日志和存盘只在 rank 0 做。
is_main = rank == 0
if is_main and step % args.log_every == 0:
print(...)
不然四个进程会往屏幕上刷四份一样的日志,checkpoint 也会被互相覆盖。
2.3 通信量算一算
00-pretrain-0.5b/parallel_math.py 的实际输出:
二、每卡每步通信量
方案 通信量 NVLink PCIe
DDP + no_sync(每步只同步一次) 1.51 GB 2.5ms 60.3ms
DDP 无 no_sync(每次累积都同步) 24.11 GB 40.2ms 964.2ms
FSDP full shard 36.16 GB 60.3ms 1446.3ms
三、通信占一步的比例(一步约 6.42 秒)
方案 NVLink PCIe
DDP + no_sync(每步只同步一次) 0.04% 0.94%
DDP 无 no_sync(每次累积都同步) 0.63% 15.02%
FSDP full shard 0.94% 22.53%
先看第一行。DDP 每步要 all-reduce 一遍全部梯度,bf16 梯度是 1.004 GB,乘上 ring 的系数 ,得到 1.51 GB。
在 NVLink 上这只要 2.5 ms,占一步(6.42 秒)的 0.04%,基本免费。在 PCIe 上是 60 ms,占 0.94%,也还好。
结论是:0.5B 模型配 16 的梯度累积,DDP 的通信开销几乎可以忽略。 这个结论有点反直觉,很多人以为多卡一定被通信拖累。原因下一节说。
2.4 no_sync:最重要的一个优化
注意上表第二行,不用 no_sync 的话通信量是 24.11 GB,整整 16 倍。
原因在于 DDP 的默认行为是「每次 backward() 都触发 all-reduce」。而我们有 16 次梯度累积,就是 16 次 backward(),于是同步了 16 遍。但前 15 次的梯度还没累加完,同步它们毫无意义。
no_sync() 这个上下文管理器就是用来关掉中间那些同步的:
for micro in range(args.grad_accum):
x, y = get_batch(train_data, args.micro_batch, args.seq_len, device)
last = (micro == args.grad_accum - 1)
sync_ctx = nullcontext() if last else model.no_sync()
with sync_ctx:
with ctx:
_, loss = model(x, y)
loss = loss / args.grad_accum
loss.backward()
只有最后一次累积走正常路径触发 all-reduce,前 15 次都在 no_sync() 里,梯度只在本地累加。
这就解释了 2.3 节那个反直觉的结论:梯度累积把通信摊薄了。累积次数越多,同一次通信服务的计算越多,通信占比越低。反过来说,如果 grad_accum=1,通信占比会直接乘 16。
这行代码少写不会报错,loss 曲线也完全正常,只是白白多花十几倍的通信。在 PCIe 机器上这个差别是 0.94% 对 15.02%,很可观。
2.5 用 torchrun 启动
torchrun --standalone --nproc_per_node=4 train_ddp.py --mode ddp
--nproc_per_node=4 表示起 4 个进程,一个进程一张卡。--standalone 是单机模式,多机要改成指定 master 地址和端口。
不要用 python train_ddp.py,那样只会起一个进程,环境变量也不会注入。
三、显存不够的时候:ZeRO 与 FSDP
3.1 DDP 浪费在哪
DDP 每张卡都存一份完整的参数、梯度、优化器状态。四张卡就是四份一模一样的优化器状态。
02 篇算过每参数 16 字节,其中 12 字节是优化器状态(fp32 master weights 加 Adam 的两个动量)。这 12 字节在四张卡上完全重复,是纯浪费。
ZeRO(Zero Redundancy Optimizer)要解决的就是这个。
3.2 ZeRO 的三个阶段
按切分程度分三级,一级比一级省:
| 阶段 | 切什么 | 每参数字节(4 卡) |
|---|---|---|
| ZeRO-0(就是 DDP) | 什么都不切 | 16 |
| ZeRO-1 | 切优化器状态 | 7 |
| ZeRO-2 | 再切梯度 | 5.5 |
| ZeRO-3 | 再切参数本身 | 4 |
ZeRO-3 就是 PyTorch FSDP 的 FULL_SHARD 模式。
四个阶段各自切了什么,画出来:
ZeRO-3 的思路:平时每张卡只存 1/N 的参数,要用到某一层的时候临时 all-gather 出完整参数,算完立刻扔掉。这样显存里任何时刻只有一层的完整参数,其余都是分片。
3.3 显存账
00-pretrain-0.5b/parallel_math.py 的输出:
一、每卡静态显存(不含激活值)
方案 字节/参数 显存
DDP / ZeRO-0(什么都不切) 16.00 8.04 GB
ZeRO-1(切优化器状态) 7.00 3.52 GB
ZeRO-2(再切梯度) 5.50 2.76 GB
ZeRO-3 / FSDP full shard(全切) 4.00 2.01 GB
从 8.04 GB 降到 2.01 GB,省了四倍,正好等于卡数。
但对我们这个模型,这个收益毫无意义,因为 8.04 GB 本来就装得下,省到 2 GB 只是把用不上的显存空出来。
3.4 FSDP 的代价
省显存不是白省的。看 2.3 节那张通信表的第三行:FSDP 是 36.16 GB,比 DDP 加 no_sync 多 24 倍。
为什么这么多。FSDP 每次前向要 all-gather 一遍参数,反向还要再 all-gather 一遍(因为前向算完就扔了),梯度还要 reduce-scatter。三趟,而且每个 micro batch 都要走一遍,没法像 DDP 那样靠 no_sync 摊薄。
把 DDP 和 FSDP 的一步画在一起,这三趟就很显眼了:

docs/images/megatron_fsdp/DDP_vs_FSDP.png,BSD-3。看下半部分那两个标着 N LAYERS 的灰框 —— all-gather 不是整个模型来一次,而是每一层各来一次。这也是为什么 FSDP 的通信量不只是「多几趟」,而是跟层数挂钩。